Skip to content

Raise Shampoo preconditioners to -1/(2k) per the paper - #548

Open
shaneraphel wants to merge 2 commits into
jettify:masterfrom
shaneraphel:shampoo-paper-exponent
Open

shaneraphel wants to merge 2 commits into
jettify:masterfrom
shaneraphel:shampoo-paper-exponent

Conversation

@shaneraphel

@shaneraphel shaneraphel commented Sep 25, 2026 •

Copy link
Copy Markdown

Summary

Gupta, Koren, Singer 2018 put each mode factor of an order-k tensor to the power -1/(2k): -1/4 on each side for matrices, -1/2 for vectors. The code used -1/k, twice the published exponent on every mode.

On a 20×6 least-squares toy at the same learning rate, the paper exponent reaches 0.523 final MSE against 0.817 for the old exponent.

Fixes #502.

Test plan

  • pytest tests/test_shampoo_exponent.py: X^4 = M^{-1} for a matrix factor, X^2 = M^{-1} for a vector factor, and the toy comparison above
  • pytest tests/test_optimizer_with_nn.py -k "Shampoo or Adahessian": 2 passed
  • tests/test_optimizer.py -k Shampoo has 7 failures that exist without this change (state-dict precision on this torch version)

Prepared with an AI assistant. I reviewed the diff and ran the tests on CPU.

get_trace averaged the Hessian diagonal over the spatial dims only for
4D kernels. Conv1d (3D) and Conv3d (5D) weights left tmp_output
unbound. Average over dims 2..ndim-1 instead; 4D is unchanged.
Each mode factor of an order-k tensor enters to the power -1/(2k):
-1/4 per side for matrices. The code used -1/k, twice the published
exponent. The 7 pre-existing test_optimizer Shampoo failures are a
state-dict precision issue on this torch version and fail identically
without this change.

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

Implementation of Shampoo inconsistent with the paper

1 participant